KV Cache 的内存占用分析
5.2.1 不同模型的KV Cache占用对比
下表展示了主流开源模型的 KV Cache 配置和内存占用估算(假设 batch_size=1,seq_len=4096, fp16精度):
| 模型 | 层数 (L) | 头数 (H) | 头维度 (K) | 隐藏维度 | KV Cache/Token (KB) | 4K序列 (MB) | 128K序列 (GB) |
|---|---|---|---|---|---|---|---|
| LLaMA-2-7B | 32 | 32 | 128 | 4096 | 16 | 64 | 2.0 |
| LLaMA-2-13B | 40 | 40 | 128 | 5120 | 25 | 100 | 3.1 |
| LLaMA-2-70B | 80 | 64 | 128 | 8192 | 80 | 320 | 10.0 |
| LLaMA-3-8B | 32 | 32 | 128 | 4096 | 16 | 64 | 2.0 |
| LLaMA-3-70B | 80 | 64 | 128 | 8192 | 80 | 320 | 10.0 |
| Mistral-7B | 32 | 32 | 128 | 4096 | 16 | 64 | 2.0 |
| Mixtral-8x7B | 32 | 32 | 128 | 4096 | 16 | 64 | 2.0 |
| Qwen-7B | 32 | 32 | 128 | 4096 | 16 | 64 | 2.0 |
| Qwen-72B | 80 | 64 | 128 | 8192 | 80 | 320 | 10.0 |
5.2.2 LLaMA-13B 详细计算示例
以LLaMA-2-13B为例,详细计算KV Cache的内存占用:
模型配置:
- 层数 \(L = 40\)
- 注意力头数 \(H = 40\)
- 每个头的维度 \(K = 128\)
- 隐藏层维度 = 5120
单token的KV Cache计算:
\[\begin{aligned} \text{KV Cache per token} \\ &= 2 \times L \times H \times K \times D \\ &= 2 \times 40 \times 40 \times 128 \times 2 \text{ (FP16)} \\ &= 2 \times 40 \times 40\times 128 \times 2 \\ &= 819,200 \text{ bytes} \\ &= 800 \text{ KB} \\ &\approx 0.78 \text{MB} \end{aligned}\]
不同序列长度的KV Cache占用
| 序列长度 | KV Cache大小 (FP16) | KV Cache大小 (INT8) | KV Cache大小 (INT4) |
|---|---|---|---|
| 1K (1024) | 800 MB | 400 MB | 200 MB |
| 4K (4096) | 3.2 GB | 1.6 GB | 800 MB |
| 8K (8192) | 6.4 GB | 3.2 GB | 1.6 GB |
| 32K (32768) | 25.6 GB | 12.8 GB | 6.4 GB |
| 128K (131072) | 102.4 GB | 51.2 GB | 25.6 GB |
| 1M (1048576) | 819.2 GB | 409.6 GB | 204.8 GB |
批处理的影响:
对于batch_size=16,序列长度=4096:
\(\text{KV Cache} = 16 \times 3.2 \text{ GB} = 51.2 \text{ GB}\)
这已经超过了单张A100 GPU的80GB显存容量的一半!
5.2.3 长上下文的影响
随着上下文窗口的扩大,KV Cache的内存占用呈线性增长:
latex
KV Cache Memory vs Sequence Length (LLaMA-70B, batch=1)
┌────────────────────────────────────────────────────────────┐
│ │
│ 12 GB ┤ ████ │
│ 10 GB ┤ ████ ████ │
│ 8 GB ┤ ████ ████ ████ │
│ 6 GB ┤ ████ ████ ████ ████ │
│ 4 GB ┤ ████ ████ ████ ████ ████ │
│ 2 GB ┤ ████ ████ ████ ████ ████ ████ │
│ 0 GB ┼────┬────┬────┬────┬────┬────┬────┬────┬──── │
│ 4K 8K 16K 32K 64K 128K 256K 512K 1M │
│ │
└────────────────────────────────────────────────────────────┘长上下文带来的挑战:
- 显存瓶颈:128K上下文需要10GB+的KV Cache,限制了batch size
- 计算效率:注意力计算复杂度为 \(O(n^2)\),长序列计算开销大
- 推理延迟:需要加载大量KV Cache,增加首token延迟
- 内存碎片:动态序列长度导致内存碎片问题
5.2.4 批大小的影响
| Batch Size | 4K序列 (LLaMA-70B) | 32K序列 (LLaMA-70B) | A100-80G占比 |
|---|---|---|---|
| 1 | 320 MB | 2.56 GB | 3.2% |
| 4 | 1.28 GB | 10.24 GB | 12.8% |
| 8 | 2.56 GB | 20.48 GB | 25.6% |
| 16 | 5.12 GB | 40.96 GB | 51.2% |
| 32 | 10.24 GB | 81.92 GB | 102.4% ⚠️ |
关键观察:
- 当batch_size=32且序列长度为32K时,KV Cache alone就超过了A100-80G的显存容量
- 这还不包括模型权重(约140GB for FP16)和激活值
- 实际生产环境中,batch_size往往受限于KV Cache而非模型权重